iT邦幫忙

2026 iThome 鐵人賽

DAY 8
0
Software Development

GPU 效能優化實戰:30 天從 Kernel 到 Profiling (重賽版)系列 第 8

Day 8|INT8 GEMM:先把答案算對(重賽版)

  • 分享至 

  • xImage
  •  

前兩天我們做了兩件事:

Day 6:把 Phase 22 拆成三段 pipeline
Day 7:把 CRT 表和 reconstruct kernel 講清楚

今天進到真正會乘矩陣的地方。

Day 8 :

7 張 INT8 餘數矩陣,最笨的乘法怎麼寫才是對的?

因為只要餘數矩陣錯,或 naive INT8 GEMM 錯,後面換成 Tensor Core 只會把錯誤算得更快。

這幾天都沒時間審稿 只能給一堆ai slot 希望之後可以有時間慢慢看 慢慢改
https://i.imgur.com/4UOZl0h.jpg

阿不過如果只想看完整的fp64或是更高 可能還要再等幾天 我會再把這部分放上來
等這幾天比較簡單(7 prime number的例子) 結束 就會說 更長的差異
先醬 peace out~
(另外我們還是以fp64作為i/o只是我們精度只有 34bit 在這個案例中 更完整精度的會再之後 3Q)
(但其實道理是差不多的 只是 mod數比較少 還有實作比較雛形 以及沒有太完整的介紹到使用的硬體以及為什麼要這樣用 bla bla bla )


先把 double 變成 high/low 整數薄片

Phase 22 的輸入是 FP64:

A: double[m, k]
B: double[k, n]

Tensor Core 不能直接吃 double,所以第一步是 Ozaki split。

程式用兩個 scale:

sh = 2^17
sl = 2^34

對每個元素 v,先取 high:

hi = round(v * sh)

再把 high 拿掉,取 low:

lo = round((v - hi / sh) * sl)

這裡的 hilo 是整數。最後重建 double 時才乘回:

hi 對應 2^-17
lo 對應 2^-34

所以兩個矩陣相乘會有四個 pass:

A_hi × B_hi
A_hi × B_lo
A_lo × B_hi
A_lo × B_lo

對應縮放:

2^-34, 2^-51, 2^-51, 2^-68

今天先拿其中一個 pass 想就好。假設我們要算 A_hi × B_hi


precompute_modulo:把薄片變成 7 份餘數

high/low 整數還是太大。

例如 v[-1, 1] 附近時:

round(v × 2^17) 可能接近 ±131072

這超過 INT8。

所以對每個質數取餘數:

const int pr[7] = {127, 113, 109, 107, 103, 101, 97};

for (int p = 0; p < 7; ++p) {
    int32_t rem = iv % pr[p];
    if (rem < 0) rem += pr[p];
    d[p * padded_size + out_off] = (int8_t)rem;
}

這就是 precompute_modulo_hl_Aprecompute_modulo_hl_B 在做的事。

它們一次輸出兩份:

d_h:high slice 的餘數
d_l:low  slice 的餘數

而且每份都有 7 個質數。

所以 A 會變成:

A_hi_mod[7]
A_lo_mod[7]

B 也一樣:

B_hi_mod[7]
B_lo_mod[7]

這些才是真正會被 INT8 GEMM 讀進去的矩陣。


A 和 B 的 layout 不一樣

Day 3 的 naive GEMM 可以直接把矩陣想成 row-major:

A[row, k]
B[k, col]

但 Phase 22 後面要餵 mma.sync.m16n8k32.s8,所以 precompute_modulo 不是單純照 row-major 寫出來。

A 被排成:

128 × 32 tile

B 被排成:

64 × 32 tile

原因是後面的 C tile 是:

128 × 64

K 方向每次吃:

32

也就是:

A tile: 128 rows × 32 K
B tile:  64 cols × 32 K
C tile: 128 rows × 64 cols

這個 layout 不是為了讓 Day 8 naive 版本舒服,而是為了 Day 9 的 Tensor Core 版本。

Day 8 如果只是寫一顆教學用 naive GEMM,可以先用正常 row-major 的 A8、B8 對答案。等答案對了,再把同一件事搬到 tiled layout。

工程上就是兩層:

正確性層:
    A8[row, k], B8[k, col]

效能層:
    A8 tiled 128×32, B8 tiled 64×32

不要一開始就把兩層混在一起。
(至於為什麼是這個數字 這個其實要做實驗跟評測 不是說隨便一個數字組合就可以了)

最笨的 INT8 GEMM

最小版本跟 Day 3 的 naive GEMM 幾乎一樣,只是輸出先不要急著變成 double。

對 CRT 來說,單一質數的 GEMM 產物應該先是一張 int32 accumulator:

__global__ void naive_int8_residue_gemm(
    const int8_t* A8,
    const int8_t* B8,
    int32_t* Cp,
    int m, int n, int k)
{
    int row = blockIdx.y * blockDim.y + threadIdx.y;
    int col = blockIdx.x * blockDim.x + threadIdx.x;

    if (row < m && col < n) {
        int32_t sum = 0;
        for (int kk = 0; kk < k; ++kk) {
            sum += (int32_t)A8[row * k + kk]
                 * (int32_t)B8[kk * n + col];
        }
        Cp[row * n + col] = sum;
    }
}

一個 thread 算一個 C[i,j]

這顆 kernel 有三個特點:

  1. 很好懂。
  2. 很好對答案。
  3. 很慢。

慢的原因也很明顯:每個 thread 自己跑完整個 K 迴圈,沒有 tiling,沒有 shared memory reuse,也沒有 Tensor Core。

但它有一個很重要的用途:確認 A8 @ B8 這件事本身是對的。

等 7 個質數的 Cp 都算完,下一步才是:

v_p = Cp mod p
x   = CRT(v_127, v_113, ..., v_97)
C  += x × scale

也就是說,scale 不屬於單一質數的 naive GEMM。它屬於 CRT 重建後的 Ozaki pass。


對答案要對「餘數」,不是直接對 double

下面的解釋是拿 CPU 的初級實作跟GPU的版本對答案
為什麼現在需要這個?
因為之後如果你要去對硬體位置來寫跟計算 會容易出錯
所以這邊用CPU的例子 做雙向除錯

今天最容易犯的錯,是一開始就拿最後的 double C 去比 cublasDgemm

那樣錯了也很難查。

Day 8 應該分層對:

第一層:單一元素的餘數

拿一個 double:

v = -0.123456

CPU 算:

hi = round(v × 2^17)
lo = round((v - hi / 2^17) × 2^34)

再算:

hi mod 127
hi mod 113
...

GPU 的 precompute_modulo_hl_A/B 輸出要一樣。

第二層:單一質數的 INT8 GEMM

固定 p = 127

CPU 算:

C_127 = A_127 @ B_127
C_127 = C_127 mod 127

GPU naive INT8 GEMM 的 sum 再 mod 127,要對得上。

第三層:7 個質數 CRT 重建

把 7 個 C_p 丟給 Day 7 的 reconstruct:

x = Σ (C_p mod p) × f_p  mod M

看它能不能拼回原本那個整數內積。

第四層:乘回 Ozaki scale

最後才做:

C += x × 2^-34

這時才拿 double 去跟 CPU 或 native DGEMM 比。

這樣切開以後,錯誤會很好定位:

單元素錯 → precompute_modulo
單質數 GEMM 錯 → INT8 multiply / layout
CRT 錯 → coeff / signed fold
double 錯 → scale / pass mapping

(這裡都是再說工程上大家可以用這裡大部分都在說大家可以用 CPU 和 GPU 來做 matching。因為我們在第六天其實就講過這個東西怎麼用 CPU 去算,所以其實你可以先寫一個簡單、你懂的 CPU 版本,然後再對上 GPU 的版本。

這樣你在除錯(debug)時會比較好處理,也比較方便。)


負數 % 是 Day 8 最大的 debug 點

昨天講過一次,今天會真的踩到。

C++ 的:

iv % p

如果 iv 是負數,結果也可能是負數。

但我們要塞進 int8_t 的是 [0, p)

所以一定要:

int32_t rem = iv % p;
if (rem < 0) rem += p;

同樣地,INT8 GEMM 的 accumulator 也可能是負的嗎?

在 Phase 22 的餘數表示下,輸入餘數是 0..126,所以單次乘法是非負,sum 也是非負。這條 fast path 裡可以用 uint32_t 的 magic remainder。

但如果你用別的表示法,例如把餘數中心化成 [-p/2, p/2],那 accumulator 就會有負數。這時 reconstruct 前又要折回 [0, p)

所以 Day 8 的建議是:先用最無聊的表示法。

所有餘數都存 [0, p)
所有 CRT 輸入都先保證 [0, p)

等答案穩了,再談有號表示能不能省一點東西。


為什麼 naive 版本一定慢

今天這顆 naive INT8 GEMM 不走 Tensor Core。

它用的是普通 CUDA core:

load int8
轉 int32
乘法
加到 int32 sum

每個 thread 算一個輸出,K 次迴圈全部自己做。A、B 的資料也沒有被 block 內其他 thread 好好重用。

更麻煩的是,Phase 22 不是只做一次 GEMM:

4 個 Ozaki pass × 7 個質數 = 28 次 INT8 GEMM

如果每次都 naive,速度一定炸掉。

所以 Day 8 的 naive kernel 不是最終實作,它是測試儀器。

它回答的是:

我們餵給 Tensor Core 的東西,數學上到底對不對?

Day 9 才回答:

怎麼讓同一件事真的走到 INT8 Tensor Core?

跟 Phase 22 fast path 的關係

現在 Phase 22 fast path 沒有真的呼叫 Day 8 這顆 naive kernel。

它做的是:

precompute_modulo_hl_A/B
    ↓
crt_fast_kernel

crt_fast_kernel 裡面已經包含:

讀 tiled A8/B8
cp.async 搬到 shared memory
ldmatrix 讀 fragment
mma.sync 做 INT8 GEMM
CRT reconstruct
寫回 double C

Day 8 把「INT8 GEMM」單獨拿出來,是為了把正確性講清楚。

更精確地說,Day 8 是 fast path 的參考答案。

fast path 做了很多效能相關的事情:

global memory layout 改排
shared memory staging
cp.async pipeline
ldmatrix fragment load
mma.sync accumulator
CRT inline reconstruct

任何一個環節錯了,最後的 double C 都會錯。

所以開發時不要直接拿 fast path 跟 cublasDgemm 硬比。比較好的方式是讓 Day 8 的 naive path 先產生一份中間答案,然後一層一層對:

precompute_modulo output
    ↓
單一 prime 的 int32 accumulator
    ↓
accumulator mod p
    ↓
7-prime CRT reconstruct
    ↓
乘回 Ozaki scale 後的 double C

fast path 至少要能在這些切點上對到 naive path。這比最後只看 max error 有用很多。

尤其 layout 這一段最容易出錯。

Day 8 的教學版可以先假設:

A8[row, kk]
B8[kk, col]

但 fast path 讀的是 tiled layout。這代表同一個數學元素在 memory 裡的位置可能完全不同。debug 時我會先準備一個很小、很容易辨認的矩陣:

A[row, kk] = 1000 * row + kk
B[kk, col] = 1000 * kk + col

然後只看 precompute_modulo 後的某一個 prime,例如 p = 127。如果某個位置本來應該是:

(1000 * row + kk) mod 127

但 device buffer 裡讀到別的值,就先不要碰 GEMM。這通常表示 layout mapping 或 padding offset 已經錯了。

padding 也要特別看。

因為 Tensor Core path 會把 M、N、K 補到 tile 友善的大小,實際 buffer 裡可能有一部分是 padding。padding 區域必須是 0,不能留下未初始化資料。否則 naive row-major 測試可能沒事,但 tiled fast path 在 K tail 或邊界 tile 會把垃圾值乘進 accumulator。

這類 bug 的特徵是:

小尺寸剛好整除 tile 時正常
一換成非 128/64/32 倍數就爆
錯誤集中在矩陣右邊或下邊

如果看到這種 pattern,先查 padding,不要先改 CRT。

repo 裡也有獨立的 INT8 GEMM kernel:

int8_gemm_s8s8s32(...)

它是後來把 s8×s8→s32 GEMM 孤立出來的版本,使用 128×128 或 128×256 tile、cp.asyncldmatrixmma.sync。那已經比較接近 Day 9/10 的內容。

今天先不要跳到那裡。

先把最笨的 A8 @ B8 寫對。


工程筆記:真的跑 repo 時先看避坑頁

  1. 連續 execute 如果結果像被加了好幾次,先查 stream 和 memset dependency。
  2. max error 如果直接爆到幾千,不要先怪 Ozaki,先查負數 modulo、layout、shared memory race。

Day 8 的重點是「每一層都能對答案」。這些 handover 筆記就是在提醒:很多 bug 看起來像數學錯,其實是環境、layout 或同步問題。

我會把實際 debug 順序寫得更死一點。

第一步,只測 precompute_modulo_hl_A/B

輸入不要先用 random。先用固定小矩陣,手動包含:

0
1
-1
127
-127
128
-128
一個普通小數,例如 0.125
一個負小數,例如 -0.125

這樣可以一次抓三種錯:

round / scale 錯
負數 modulo 錯
int8 cast 錯

第二步,只測單一 prime 的 GEMM。

先固定 p = 127,只看:

A_127 @ B_127

這一步不要做 CRT,也不要乘回 double scale。輸出就是 int32 accumulator,然後 CPU 也用同一批 A_127B_127 算一份。兩邊 accumulator 對上後,再比較 acc % 127

如果 accumulator 對不上,問題通常在:

A/B layout
K 迴圈範圍
padding 是否清 0
row/col index mapping

如果 accumulator 對得上,但 acc % p 對不上,才去查 remainder implementation。

第三步,才測 7 個 prime 的 CRT。

這一步要刻意測負的整數內積。例如選 A 和 B,讓某個 C[i,j] 是負數。因為 CRT reconstruct 會先得到 [0, M) 裡的代表值,再靠:

if x > M/2:
    x = x - M

折回 signed integer。沒有測負內積,就等於沒有測這條 signed fold。

第四步,才測 Ozaki 四個 pass 的 scale。

四個 pass 的 scale 不一樣:

hi × hi: 2^-34
hi × lo: 2^-51
lo × hi: 2^-51
lo × lo: 2^-68

如果 CRT 的整數內積都對,但 double 結果還是錯,就要查 pass mapping 和 scale。這時候不要回頭亂改 prime 表。

最後才跑大尺寸 random matrix。

random matrix 很適合找殘留 bug,但不適合作為第一個測試。它的問題是資訊量太低:你只會知道 max error 很大,卻不知道錯在哪一層。Day 8 的做法是先把每一層都變成可以被單獨觀察的東西,讓後面的 Tensor Core 優化有一個可信的基準。


今天的產出

Day 8 結束時,應該有四個東西:

  1. precompute_modulo 可以把 high/low 轉成 7 份 int8_t 餘數。
  2. 單一元素的餘數能跟 CPU 對上。
  3. naive INT8 GEMM 能跟 CPU 的 A8 @ B8 對上。
  4. naive GEMM 的結果丟進 CRT reconstruct 後,可以回到正確的整數內積。

這時候它一定很慢。

慢是正常的。

今天的成功標準不是 ms,而是:

每一層都能對答案
錯了知道是哪一層

今天先停在這裡

今天可以帶走三句:

  1. precompute_modulo 是 double 到 INT8 世界的入口。 它把 high/low 整數薄片轉成 7 份餘數矩陣。
  2. naive INT8 GEMM 是正確性工具。 一個 thread 算一個 C[i,j],很慢,但最好查。
  3. 先對餘數,再對 double。 不要一開始就看最後誤差,否則錯在取餘數、layout、CRT、scale 都混在一起。

明天把同一個 A8 @ B8 搬上 Tensor Core:128×64 tile、K=32、cp.asyncldmatrix,最後用 mma.sync.m16n8k32.s8 真的把 INT8 乘法吃起來。

額外的 murmur ── 既然是重賽版,我們需要寫一些新的東西,所以我這裡會放一些 murmur,或者放一些派大星、海綿寶寶之類的梗圖,希望有趣一點。然後我也會多放一些心得跟想法,多寫一些字。其實我知道自己沒有設計得很好,但我想說,反正我就努力維持每天都有寫。

我的想法是,我可能有好幾天都是 AI slot,但我其中會有幾天是很用心、很認真寫的。像我前三天基本上都是很認真、用心寫的東西。在這個專案(這個 Adapted GEMM)開始之前,我基本上都偏向自己寫,而且自己用心在寫;在之後呢,當然比較多 AI slot 的部分,但也是我實際上做過的專案,我覺得很值得分享。

後面這邊也放一些 murmur 吧。我覺得今天的 murmur 可以講很多東西,像是環境或依賴性?

你看,我們在建構這個 kernel 函式的時候也是一樣,必須要有一個 source 來源是正確的(CPU),或是可以幫助我們去 debug 的。你用 AI 在寫東西的時候也是一樣,AI 講的東西不一定對,你至少要自己 review 過,要嘛就是你要有一個 formal way 讓它去產生一個 definitely true 的解,還要有一個正規表示式,讓它表達出來是正確的。這也是為什麼會有向量機,對吧?

那奇怪的是,明明是一個非常確定性的解、確定性的答案,為什麼會導致這樣的問題呢?明明是一個確定解,明明大部分都是確定解,為什麼會有這樣的問題?

那就是因為像圖靈機也不是都可以停機的,而像是我們的 GPU 算出來也不一定都是正確的,誰知道 compiler 又會作什麼妖呢?

我覺得都不知道耶。有時候有些事情,你可能忙了一輩子,或是忙了好一陣子,花了很多時間投注在上面,突然間某天早上起來,會發現感覺好沒有意義喔。

今天的 murmur 大概是這醬 祝好 明天見。


上一篇
Day 7|CRT 表:7 個質數怎麼拼回一個大整數 (重賽版)
下一篇
Day 9|INT8 GEMM:真的走上 Tensor Core (重賽版)
系列文
GPU 效能優化實戰:30 天從 Kernel 到 Profiling (重賽版)17
圖片
  熱門推薦
圖片
{{ item.channelVendor }} | {{ item.webinarstarted }} |
{{ formatDate(item.duration) }}
直播中

尚未有邦友留言

立即登入留言